#!/usr/bin/env python3
"""complaint_information.md를 기계적으로 조립하여 complaint_document.md를 생성한다."""

from __future__ import annotations

import html
import itertools
import json
import re
import sys
import unicodedata
from collections import OrderedDict
from dataclasses import dataclass
from typing import Any

import httpx


# =============================================================================
# MCP / constants
# =============================================================================
LOCALDOCS_URL = "http://mcp-localdocs:8012/mcp"
MCP_HEADERS = {
    "Content-Type": "application/json",
    "Accept": "application/json, text/event-stream",
}
CLIENT = httpx.Client(timeout=60)
MSG_ID_COUNTER = itertools.count(100)

SOURCE_DOC_NAME = unicodedata.normalize("NFC", "complaint_information.md")
OUTPUT_DOC_NAME = "complaint_document.md"


def next_msg_id() -> int:
    return next(MSG_ID_COUNTER)


def _init_mcp_session() -> None:
    response = CLIENT.post(
        LOCALDOCS_URL,
        json={
            "jsonrpc": "2.0",
            "id": 1,
            "method": "initialize",
            "params": {
                "protocolVersion": "2025-03-26",
                "capabilities": {},
                "clientInfo": {"name": "stage5-complaint-doc", "version": "1.0"},
            },
        },
        headers=MCP_HEADERS,
    )
    response.raise_for_status()

    session_id = response.headers.get("mcp-session-id")
    if session_id:
        MCP_HEADERS["mcp-session-id"] = session_id

    CLIENT.post(
        LOCALDOCS_URL,
        json={"jsonrpc": "2.0", "method": "notifications/initialized"},
        headers=MCP_HEADERS,
    ).raise_for_status()


def _parse_mcp(text: str) -> dict[str, Any] | None:
    for line in text.strip().split("\n"):
        if line.startswith("data: "):
            try:
                return json.loads(line[6:])
            except json.JSONDecodeError:
                return None

    try:
        return json.loads(text)
    except json.JSONDecodeError:
        return None


def _call_tool(name: str, arguments: dict[str, Any], msg_id: int) -> dict[str, Any]:
    response = CLIENT.post(
        LOCALDOCS_URL,
        json={
            "jsonrpc": "2.0",
            "id": msg_id,
            "method": "tools/call",
            "params": {"name": name, "arguments": arguments},
        },
        headers=MCP_HEADERS,
    )
    response.raise_for_status()

    parsed = _parse_mcp(response.text)
    if not parsed or "result" not in parsed:
        raise RuntimeError(f"MCP {name} failed: {response.text[:300]}")
    return parsed


def _unwrap_text(parsed: dict[str, Any], doc_name: str) -> str:
    text = (parsed["result"].get("content") or [{}])[0].get("text", "")
    if not text:
        raise RuntimeError(f"Empty response for {doc_name}")

    try:
        outer = json.loads(text)
    except json.JSONDecodeError:
        return text

    if isinstance(outer, dict) and "results" in outer:
        results = outer.get("results") or []
        if not results:
            raise RuntimeError(f"No results for {doc_name}")

        inner = results[0].get("content") or results[0].get("text") or ""
        if not inner:
            raise RuntimeError(f"Empty content for {doc_name}")
        return inner if isinstance(inner, str) else json.dumps(inner, ensure_ascii=False)

    return json.dumps(outer, ensure_ascii=False) if not isinstance(outer, str) else outer


def read_text(doc_name: str) -> str:
    parsed = _call_tool("read_docs", {"doc_names": [doc_name]}, next_msg_id())
    return _unwrap_text(parsed, doc_name)


def write_doc(path: str, content: str) -> None:
    _call_tool("write_file", {"path": path, "content": content}, next_msg_id())


# =============================================================================
# Parsing
# =============================================================================
OVERVIEW_HEADER = "## 1. 청구권 개요"
DETAIL_HEADER = "## 2. 개별 청구권 상세 정보"
EVIDENCE_HEADER = "## 3. 서증 목록"
NEXT_TOP_LEVEL_HEADER_PATTERN = re.compile(r"^##\s+\d+\.", re.MULTILINE)
DETAIL_BLOCK_PATTERN = re.compile(
    r'(?ms)^####\s+(C-\d{3}):\s*(.*?)\n\n(<div class="full-width-table-wrap".*?</div>)\s*'
)
TABLE_ROW_PATTERN = re.compile(r"(?is)<tr\b[^>]*>(.*?)</tr>")
TABLE_CELL_PATTERN = re.compile(r"(?is)<t[dh]\b[^>]*>(.*?)</t[dh]>")
BR_PATTERN = re.compile(r"(?i)<br\s*/?>")
TAG_PATTERN = re.compile(r"(?is)<[^>]+>")
TOP_NUMBER_PATTERN = re.compile(r"^\s*(\d+)\.\s*(.*)$")
SUBITEM_PATTERN = re.compile(r"^\s*([가나다라마바사아자차카타파하])\.\s*(.*)$")
PLAINTIFF_REF_PATTERN = re.compile(r"(?:(주위적으로,|예비적으로,)\s*)?피고는\s+원고(에게|에\s+대하여)")
SUBITEM_LABELS = ["가.", "나.", "다.", "라.", "마.", "바.", "사.", "아.", "자.", "차."]


@dataclass
class Claim:
    claim_id: str
    claim_title: str
    plaintiff: str
    defendant: str
    relief_text: str
    cause_text: str


def claim_sort_key(claim_id: str) -> int:
    return int(claim_id.split("-")[1])


def find_section_bounds(text: str, start_header: str, end_header: str | None = None) -> tuple[int, int]:
    start = text.find(start_header)
    if start == -1:
        raise RuntimeError(f"헤더를 찾을 수 없습니다: {start_header}")

    section_start = start + len(start_header)
    if end_header:
        end = text.find(end_header, section_start)
        if end == -1:
            raise RuntimeError(f"종료 헤더를 찾을 수 없습니다: {end_header}")
        return section_start, end

    match = NEXT_TOP_LEVEL_HEADER_PATTERN.search(text, section_start)
    if not match:
        return section_start, len(text)
    return section_start, match.start()


def html_to_text(fragment: str) -> str:
    fragment = BR_PATTERN.sub("\n", fragment)
    fragment = TAG_PATTERN.sub("", fragment)
    fragment = html.unescape(fragment)
    fragment = fragment.replace("\xa0", " ")
    lines = [re.sub(r"[ \t]+", " ", line).strip() for line in fragment.splitlines()]
    return "\n".join(line for line in lines if line)


def parse_overview_rows(source_text: str) -> list[list[str]]:
    start, end = find_section_bounds(source_text, OVERVIEW_HEADER, DETAIL_HEADER)
    section_text = source_text[start:end]
    rows: list[list[str]] = []

    for row_html in TABLE_ROW_PATTERN.findall(section_text):
        cells = TABLE_CELL_PATTERN.findall(row_html)
        if len(cells) != 5:
            continue
        values = [html_to_text(cell) for cell in cells]
        if values[0] == "청구 ID":
            continue
        rows.append(values)

    if not rows:
        raise RuntimeError("청구권 개요 표 행을 찾을 수 없습니다.")
    return rows


def parse_detail_fields(table_html: str) -> dict[str, str]:
    fields: dict[str, str] = {}
    for row_html in TABLE_ROW_PATTERN.findall(table_html):
        cells = TABLE_CELL_PATTERN.findall(row_html)
        if len(cells) < 2:
            continue
        key = html_to_text(cells[0])
        value = html_to_text(cells[1])
        if key and key != "항목":
            fields[key] = value
    return fields


def parse_claims(source_text: str) -> list[Claim]:
    overview_rows = parse_overview_rows(source_text)
    overview_map = {
        row[0]: {
            "claim_title": row[1],
            "plaintiff": row[3],
            "defendant": row[4],
        }
        for row in overview_rows
    }

    start, end = find_section_bounds(source_text, DETAIL_HEADER, EVIDENCE_HEADER)
    detail_section = source_text[start:end]
    claims: list[Claim] = []

    for claim_id, header_title, table_html in DETAIL_BLOCK_PATTERN.findall(detail_section):
        if claim_id not in overview_map:
            raise RuntimeError(f"청구권 개요에 없는 claim_id입니다: {claim_id}")

        fields = parse_detail_fields(table_html)
        if "청구취지" not in fields or "청구원인" not in fields:
            raise RuntimeError(f"{claim_id} 상세표에서 청구취지 또는 청구원인을 찾을 수 없습니다.")

        overview = overview_map[claim_id]
        claims.append(
            Claim(
                claim_id=claim_id,
                claim_title=overview["claim_title"] or header_title,
                plaintiff=overview["plaintiff"],
                defendant=overview["defendant"],
                relief_text=fields["청구취지"],
                cause_text=fields["청구원인"],
            )
        )

    if not claims:
        raise RuntimeError("개별 청구권 상세 정보 블록을 찾을 수 없습니다.")

    claims.sort(key=lambda claim: claim_sort_key(claim.claim_id))
    return claims


def extract_evidence_lines(source_text: str) -> list[str]:
    start, end = find_section_bounds(source_text, EVIDENCE_HEADER)
    section_text = source_text[start:end]
    return [line.rstrip() for line in section_text.splitlines() if line.strip()]


# =============================================================================
# Relief / cause assembly
# =============================================================================
def has_final_consonant(text: str) -> bool:
    for char in reversed(text.strip()):
        if "가" <= char <= "힣":
            return (ord(char) - 0xAC00) % 28 != 0
    return False


def topic_particle(text: str) -> str:
    return "은" if has_final_consonant(text) else "는"


def strip_outline_prefix(line: str) -> str:
    current = line.strip()
    while True:
        number_match = TOP_NUMBER_PATTERN.match(current)
        if number_match:
            current = number_match.group(2).strip()
            continue

        subitem_match = SUBITEM_PATTERN.match(current)
        if subitem_match:
            current = subitem_match.group(2).strip()
            continue

        break
    return current


def is_cost_clause(line: str) -> bool:
    normalized = re.sub(r"\s+", "", strip_outline_prefix(line))
    return "소송비용" in normalized and "부담" in normalized


def is_provisional_execution_clause(line: str) -> bool:
    normalized = re.sub(r"\s+", "", strip_outline_prefix(line))
    return "가집행" in normalized


def split_mode_prefix(content: str) -> tuple[str, str] | None:
    stripped = content.strip()
    for mode in ("주위적으로", "예비적으로"):
        if stripped in {mode, f"{mode},"}:
            return f"{mode},", ""

        prefix = f"{mode},"
        if stripped.startswith(prefix):
            return prefix, stripped[len(prefix) :].strip()

        if stripped.startswith(f"{mode} "):
            return prefix, stripped[len(mode) :].strip(" ,")

    return None


def build_substantive_relief_text(relief_text: str) -> str:
    lines = [line.strip() for line in relief_text.splitlines() if line.strip()]
    substantive_lines: list[str] = []

    for line in lines:
        bare = strip_outline_prefix(line)
        if is_cost_clause(bare) or is_provisional_execution_clause(bare):
            continue
        substantive_lines.append(line)

    phrases: list[str] = []
    pending_mode: str | None = None

    for line in substantive_lines:
        content = strip_outline_prefix(line)
        mode_info = split_mode_prefix(content)

        if mode_info and not mode_info[1]:
            pending_mode = mode_info[0]
            continue

        if pending_mode:
            phrases.append(f"{pending_mode} {content}".strip())
            pending_mode = None
            continue

        if mode_info:
            mode_label, remainder = mode_info
            phrases.append(f"{mode_label} {remainder}".strip())
        else:
            phrases.append(content)

    if pending_mode:
        phrases.append(pending_mode.rstrip(","))

    return " ".join(phrases).strip()


def inject_plaintiff_name(text: str, plaintiff: str) -> str:
    def replace(match: re.Match[str]) -> str:
        prefix = f"{match.group(1)} " if match.group(1) else ""
        postposition = match.group(2).replace(" ", "")
        return f"{prefix}원고 {plaintiff}{postposition}"

    updated = PLAINTIFF_REF_PATTERN.sub(replace, text, count=1)
    if plaintiff in updated:
        return updated
    return f"원고 {plaintiff}에 대하여, {updated}"


def build_relief_groups(claims: list[Claim]) -> OrderedDict[str, list[Claim]]:
    groups: OrderedDict[str, list[Claim]] = OrderedDict()
    for claim in claims:
        groups.setdefault(claim.defendant, []).append(claim)
    return groups


def render_relief_section(claims: list[Claim]) -> tuple[list[str], dict[str, str], int]:
    groups = build_relief_groups(claims)
    lines: list[str] = []
    label_map: dict[str, str] = {}

    for group_index, (defendant, grouped_claims) in enumerate(groups.items(), start=1):
        group_label = f"{group_index}."
        if len(grouped_claims) == 1:
            claim = grouped_claims[0]
            label_map[claim.claim_id] = group_label
            relief_text = build_substantive_relief_text(claim.relief_text)
            relief_text = inject_plaintiff_name(relief_text, claim.plaintiff)
            lines.append(f"{group_label} 피고 {defendant}{topic_particle(defendant)} {relief_text}")
            lines.append("")
            continue

        lines.append(f"{group_label} 피고 {defendant}{topic_particle(defendant)}")
        for sub_label, claim in zip(SUBITEM_LABELS, grouped_claims):
            label_map[claim.claim_id] = f"{group_label} {sub_label}"
            relief_text = build_substantive_relief_text(claim.relief_text)
            relief_text = inject_plaintiff_name(relief_text, claim.plaintiff)
            lines.append(f"  {sub_label} {relief_text}")
        lines.append("")

    common_cost_line = (
        "소송비용은 피고가 부담한다."
        if len(groups) == 1
        else "소송비용은 피고들이 부담한다."
    )
    lines.append(f"{len(groups) + 1}. {common_cost_line}")
    lines.append(f"{len(groups) + 2}. 제1항은 가집행할 수 있다.")

    return lines, label_map, len(groups)


def render_cause_section(claims: list[Claim], label_map: dict[str, str]) -> list[str]:
    lines: list[str] = []

    for claim in claims:
        label = label_map[claim.claim_id]
        lines.append(f"{label} 원고 {claim.plaintiff}의 피고 {claim.defendant}에 대한 {claim.claim_title}")
        cause_lines = [line.strip() for line in claim.cause_text.splitlines() if line.strip()]
        lines.extend(cause_lines)
        lines.append("")

    if lines and not lines[-1]:
        lines.pop()
    return lines


def render_party_section(title: str, parties: list[str]) -> list[str]:
    lines: list[str] = []
    for index, party_name in enumerate(parties, start=1):
        lines.extend(
            [
                f"## {title} {index}",
                f"- 성명: {party_name}",
                "- 주민등록번호:",
                "- 주소:",
                "- 연락가능한 전화번호:",
                "",
            ]
        )
    return lines


def unique_in_order(values: list[str]) -> list[str]:
    seen: set[str] = set()
    ordered: list[str] = []
    for value in values:
        if value not in seen:
            seen.add(value)
            ordered.append(value)
    return ordered


def assemble_document(source_text: str) -> str:
    overview_rows = parse_overview_rows(source_text)
    plaintiffs = unique_in_order([row[3] for row in overview_rows])
    defendants = unique_in_order([row[4] for row in overview_rows])
    claims = parse_claims(source_text)
    evidence_lines = extract_evidence_lines(source_text)

    relief_lines, label_map, defendant_count = render_relief_section(claims)
    cause_lines = render_cause_section(claims, label_map)

    lines: list[str] = [
        "# 소 장",
        "",
        "## 사건명",
        "",
    ]
    lines.extend(render_party_section("원고", plaintiffs))
    lines.extend(render_party_section("피고", defendants))
    lines.extend(
        [
            "## 청 구 취 지",
            *relief_lines,
            "",
            "## 청 구 원 인",
            *cause_lines,
            "",
            "## 입 증 방 법",
            *evidence_lines,
            "",
            "## 첨 부 서 류",
            f"1. 위 증명방법 각 {defendant_count + 1}통",
            "2. 영수필 확인서",
            "3. 송달료 납부서",
            "4. 소송위임장",
            "5. 소장 부본",
            "6. ...",
            "",
        ]
    )
    return "\n".join(lines)


# =============================================================================
# Main
# =============================================================================
def main() -> None:
    _init_mcp_session()

    source_text = read_text(SOURCE_DOC_NAME)
    print(f"[stage5] read {SOURCE_DOC_NAME} ({len(source_text)} chars)", file=sys.stderr)

    output_text = assemble_document(source_text)
    write_doc(OUTPUT_DOC_NAME, output_text)

    claims = parse_claims(source_text)
    defendant_count = len(unique_in_order([claim.defendant for claim in claims]))
    print(
        f"[stage5] saved: {OUTPUT_DOC_NAME} "
        f"(claims={len(claims)}, defendants={defendant_count})",
        file=sys.stderr,
    )
    print(
        json.dumps(
            {
                "status": "ok",
                "claims": len(claims),
                "defendants": defendant_count,
            }
        )
    )


if __name__ == "__main__":
    main()
